Skip to main content

03 - 预训练主循环

这一篇产出一个 train.py。跑起来 loss 会从 10.37 开始往下掉,能存 checkpoint,能断点续训,屏幕上能看到 MFU。

刻意不用 Trainer,不用 Accelerate。整个循环三十行以内看得完,每一行在干什么都说得清。

前置:01 篇产出的 train.binval.bin,02 篇产出的 model.py。这一篇还是单卡,多卡放到 04 篇。

零、开始之前:训练循环在循环什么

0.1 五个动作

训练就是把下面五步重复几千次,没有别的:

训练就是把这五步重复 9,537 次,没有别的① 取一批数据x, y第一节 · memmap 切片② 前向算 lossmodel(x, y)02 篇已经写完了③ 反向算梯度loss.backward()第二、三节④ 更新参数optimizer.step()第四、五、六节⑤ 梯度清零zero_grad()第三节回到第 ①,重复 9,537 步(4 卡配置下)这个循环体本身三十行就能写完。这一篇剩下的全部篇幅,都花在让这三十行连续跑满 17 小时不出事上。灰色两格已经由前两篇交付:01 篇产出 train.bin,02 篇产出 model.py。红色三格是本篇要写的。注意第 ⑤ 步不是可有可无的收尾 —— PyTorch 的梯度是累加的,忘了清零,下一步会带着上一步的梯度一起更新。
把循环画成闭环是为了强调一件事:这五步之间没有任何一步可以异步或跳过,它们严格串行。后面所有优化(混合精度、梯度累积、通信重叠)本质上都是在不改变这个顺序的前提下,让其中某一步变便宜或者藏进另一步里。

02 篇已经把第 2 步(前向和 loss)写完了,01 篇把第 1 步的数据准备好了。这一篇要补的是第 3 到 5 步,以及一堆让这个循环真的能跑完的工程细节。

0.2 为什么不是三十行就结束了

理论上确实三十行就够。但直接这么写,会遇到四个问题:

显存不够。 想要大 batch 让训练稳定,但一次前向放不下那么多样本。对策是梯度累积(第三节)。

太慢。 全用 fp32 算,速度只有一半不到。对策是混合精度(第二节)。

训不稳。 学习率从一开始就拉满,前几十步梯度会炸。对策是 warmup 加梯度裁剪(第五、六节)。

跑一半挂了。 租的实例被回收,17 小时白跑。对策是 checkpoint(第八节)。

这一篇剩下的内容基本就是在解决这四件事。

直接照着上一张图写三十行,会撞上这四件事显存不够想要大 batch 让训练稳,一次前向放不下那么多条梯度累积第三节 · grad_accum = 16太慢全用 fp32 算,速度只有一半不到bf16 混合精度第二节 · autocast训不稳学习率一开始就拉满,前几十步梯度会炸warmup + 梯度裁剪第五、六节跑一半挂了租的实例被回收,17 小时白跑checkpoint 原子写第八节这四条没有一条是可选的。少任何一条,17 小时的训练要么跑不到终点,要么跑到了但结果不能用。也可以反过来读这张图:把本篇剩下的章节全部拿掉,剩下的就是那个能跑但跑不完的三十行版本。
值得注意的是这四个问题的性质不同。前两个是效率问题,做不到只是慢;后两个是正确性和可用性问题,做不到就是全盘作废。所以真要砍工作量,第五、六、八节是最不能砍的。

0.3 这一篇要定的超参

先把结论摆着,后面每一节解释一个:

超参取值在哪一节
精度bf16 混合精度第二节
micro_batch8第三节
grad_accum16第三节
优化器AdamW,β=(0.9, 0.95)第四节
weight_decay0.1,只加在 2 维以上参数第四节
学习率3e-4,warmup 2% 后 cosine 降到 3e-5第五节
grad_clip1.0第六节

一、取数据

1.1 直接用 memmap,不用 DataLoader

01 篇 5.2 节写过 get_batch,这里补上搬到 GPU 的部分:

def get_batch(data, batch_size, seq_len, device):
ix = np.random.randint(len(data) - seq_len - 1, size=batch_size)
x = torch.from_numpy(np.stack([data[i:i + seq_len].astype(np.int64) for i in ix]))
y = torch.from_numpy(np.stack([data[i + 1:i + 1 + seq_len].astype(np.int64) for i in ix]))
if device.type == "cuda":
x = x.pin_memory().to(device, non_blocking=True)
y = y.pin_memory().to(device, non_blocking=True)
else:
x, y = x.to(device), y.to(device)
return x, y

为什么不用 DataLoader DataLoader 解决的是「样本长度不一、要 shuffle、要 collate」这些问题。而我们的数据已经是一个扁平的 token 流,取一批就是随机取几个起点做切片,没有任何需要 collate 的东西。用 DataLoader 反而要多绕一层 Dataset 抽象,还要处理多进程 worker 和 memmap 的配合,收益是负的。

pin_memorynon_blocking 是一对。 普通内存里的张量拷到 GPU,CUDA 得先把它复制到一块锁页内存再传,这个过程是同步的,会卡住 CPU。pin_memory() 直接把张量分配在锁页内存里,配合 non_blocking=True 就能让这次 H2D 拷贝异步进行,跟 GPU 上的计算重叠。

两个必须一起用。只写 non_blocking=True 而内存没锁页,PyTorch 会静默退回同步拷贝,不报错也没有加速。

1.2 随机取起点意味着不存在 epoch

传统的 epoch:把数据集从头到尾走一遍,每条恰好见一次按顺序遍历12345678走到头 = 一个 epoch 结束,「训了几轮」是个有意义的说法我们的做法:在一条 100 亿 token 的流上随机撒窗口随机取起点重叠:这段被取到两次空白:这段一次都没被取到在 100 亿 token 的尺度上,这两种偏差都可以忽略所以进度只用「走了多少步」和「已见多少 token」度量:10,000,000,000 ÷ 1,048,576 = 9,537 步就是终点,中间没有「第几轮」这回事。
这不只是叫法问题。没有 epoch 意味着没有「每轮结束时打乱数据」这个环节,也就少了一整类顺序相关的 bug;同时也意味着不能靠「跑满 N 轮」来判断训练够不够,只能盯 token 数。学习率调度按步数而不是按轮数来定,根源也在这里。

因为每一步都是随机取起点,训练过程中没有「跑完一轮数据」这个概念,也就没有 epoch。总量由步数控制,跑够 9537 步就停。

好处是不用管数据顺序,也不用 shuffle。代价是同一段文本可能被看到不止一次、也可能一次都没被完整看到。在 10B token 这个量级上,这个偏差可以忽略。

二、混合精度

2.1 为什么不全用 fp32

fp32 是 32 位浮点,精度足够,但有两个问题:占显存(每个数 4 字节),而且现代 GPU 的低精度算力远高于 fp32。A100 的 fp32 峰值只有 19.5 TFLOPS,而 BF16 走 Tensor Core 是 312 TFLOPS,差 16 倍

不用低精度,等于买了张 A100 只用了十六分之一。

2.2 fp16 和 bf16 的区别

两种 16 位浮点,位数怎么分配不一样:

格式符号位指数位尾数位能表示的范围精度
fp321823约 ±3.4e38
fp161510约 ±65504
bf16187约 ±3.4e38

位分配画出来,差别一目了然:

fp32(32 位)S指数 8 位尾数 23 位范围 ±3.4e38fp16(16 位)S指数 5 位尾数 10 位范围只到 ±65504 —— 梯度到 1e-8 量级会下溢成 0bf16(16 位)S指数 8 位尾数 7 位范围和 fp32 一样 ±3.4e38,只是精度低 ← 我们用这个指数位决定「能表示多大范围」fp16 要配 GradScaler把 loss 乘一个大系数抬进可表示区间,反向后再除回来,遇 inf 还要跳过该步。bf16 不需要动态范围与 fp32 相同,梯度再小也不下溢。Ampere(A100 / 3090)及以后一律用它。
bf16 的设计取舍很直接:牺牲尾数精度,保住和 fp32 相同的 8 位指数。训练对数值范围敏感、对精度不敏感(参数更新在 fp32 的 master weights 上做),所以这个取舍正合适。

关键在于 bf16 的指数位和 fp32 一样多,所以它能表示的数值范围跟 fp32 相同,只是精度低。fp16 反过来,精度尚可但范围窄得多,最大只到 65504。

训练时梯度经常是很小的数(1e-8 量级),fp16 直接下溢成 0,梯度就没了。所以 fp16 训练必须配一个 GradScaler:把 loss 乘上一个大系数,让梯度落进 fp16 能表示的区间,反向之后再除回来。这套机制要动态调整系数,遇到 inf 还要跳过这一步,很麻烦。

bf16 不需要 GradScaler,因为它的动态范围和 fp32 一样,梯度再小也不会下溢。代价是精度低,但训练对精度不敏感(参数更新本来就是在 fp32 的 master weights 上做的)。

结论:Ampere(A100、3090)及以后的卡一律用 bf16。只有更老的卡(V100、T4)没有 bf16 支持,才需要退回 fp16 加 GradScaler。

2.3 代码

use_bf16 = device.type == "cuda" and torch.cuda.is_bf16_supported()
ctx = (torch.autocast(device_type="cuda", dtype=torch.bfloat16)
if use_bf16 else nullcontext())

with ctx:
_, loss = model(x, y)
loss.backward()

注意 backward()with 块外面autocast 只需要覆盖前向。反向传播时 PyTorch 会自动按前向记录的 dtype 来算,把 backward() 放进去不会出错但没有意义。

2.4 autocast 不是把所有东西都转成 bf16

「混合精度」是真的混合:autocast 按算子类型逐个决定用什么精度自动转 bf16 —— 矩阵乘类nn.Linear(q/k/v/o、gate/up/down、lm_head)matmul、conv、attention 里的 QKᵀ 与 PV为什么可以:这类算子计算量最大、对精度不敏感,而 Tensor Core 上 bf16 的吞吐是 fp32 的数倍。保持 fp32 —— 归约与数值敏感类softmax、cross_entropy 的求和RMSNorm 的均方统计量、优化器更新与主权重为什么必须:归约要把成千上万个数加起来,低精度累加会一路丢有效位,误差随长度放大。同一层里,矩阵乘走 bf16、紧接着的 softmax 走 fp32,切换由 autocast 按算子查表完成,不需要手工标注每一个张量。但白名单只覆盖 PyTorch 内置算子。02 篇 RMSNorm 里那句 x.float() 就是在 autocast 之外再加的一道保险 —— 自定义算子的中间量它管不着。
把混合精度理解成「整个模型换成 bf16」是最常见的误解,也是一类数值 bug 的来源。真实情况是一张按算子分的表:算得多的走低精度,加得多的走高精度。知道这条分界线在哪,才知道自己写的算子该往哪一边靠。

autocast 内部维护了一张算子白名单。矩阵乘法、卷积这类计算密集又对精度不敏感的,走 bf16。而 softmaxlayer_norm、各种 reduction 和 loss 计算,会自动保持 fp32,因为它们涉及累加,低精度容易丢有效数字。

这也是 02 篇 RMSNorm 里为什么要手动写 x.float():那段是自己实现的,不在 autocast 的白名单机制里,得自己保证统计量用 fp32 算。

三、梯度累积

3.1 想要大 batch,但显存装不下

batch 越大,梯度的估计越准,训练越稳。大模型预训练的 batch 通常在百万 token 量级。

但一次前向能放多少样本,是硬性受显存限制的。02 篇算过静态显存 8 GB,剩下的都给激活值,micro_batch 只能开到 8 左右。8 × 2048 = 16384 token,离百万差着两个数量级。

3.2 做法:先攒够梯度再更新

梯度累积的思路很简单:跑多次前向反向,梯度自然累加在 .grad 里,攒够了再调用一次 optimizer.step()

这利用了 PyTorch 的一个默认行为:backward() 是把梯度加到 .grad 上,不是覆盖。平时每步都 zero_grad() 所以感觉不到,把 zero_grad() 挪到累积结束后,就自动实现了累积。

一个训练步内部(grad_accum = 16)micro 0前向反向micro 1前向反向micro 2前向反向micro 15前向反向裁剪step()zero_grad.grad 里累积的量攒满才更新归零等效 batch8 条 × 2048 = 16,384 tokens(一次前向)× 16 次累积 × 4 卡= 1,048,576 tokens / 步loss 必须除以 grad_accum16 个平均值直接相加 = 真实平均的 16 倍,等效学习率被放大 16 倍。不报错,训一阵才突然 NaN。
前 15 次只累加梯度,不更新参数;第 16 次之后才做一次裁剪、更新、清零。这样用 16 KB 的显存代价换到了百万 token 量级的等效 batch。04 篇会看到,这个结构还顺带把多卡通信摊薄了 16 倍。
for micro in range(args.grad_accum):
x, y = get_batch(train_data, args.micro_batch, args.seq_len, device)
with ctx:
_, loss = model(x, y)
loss = loss / args.grad_accum # 关键,见 3.3
loss.backward()

grad_norm = torch.nn.utils.clip_grad_norm_(model.parameters(), args.grad_clip)
optimizer.step()
optimizer.zero_grad(set_to_none=True)

3.3 loss 必须除以累积次数

这行 loss = loss / args.grad_accum 少了会怎样,值得说清楚。

正常的 batch 是对 N 个样本的 loss 求平均。梯度累积跑 16 次,每次的 loss 已经是各自 8 个样本的平均,直接累加得到的是 16 个平均值的,比真实平均大了 16 倍。梯度也跟着大 16 倍。

后果是等效学习率被放大了 16 倍。你以为在用 3e-4,实际是 4.8e-3,训练几十步就发散。

这个 bug 很隐蔽,因为代码不报错,loss 曲线也不是立刻爆炸,而是训一阵子突然变 NaN。

zero_grad(set_to_none=True) 是把 .grad 置为 None 而不是填 0,省一次显存写入,也省一点显存。这是现在的推荐写法。

3.4 batch 的账

把数算一遍。00-pretrain-0.5b/train_schedule.py 的实际输出:

批次账
每卡一次前向 8 条 × 2048 = 16,384 tokens
累积 16 次 262,144 tokens/卡/步
× 4 卡 1,048,576 tokens/步
总步数 10,000,000,000 / 1,048,576 = 9,537 步
warmup 190 步(2%)

每步 1,048,576 个 token,正好一百万出头,这跟主流小模型预训练的 batch 规模是一致的。总共 9537 步。

注意这里已经把 4 卡算进去了。单卡跑的话每步是 262,144 token,总步数变成 38,147 步,但总 token 数不变。卡数只影响步数,不影响看到的数据量。

四、优化器

4.1 AdamW,β2 取 0.95

用 AdamW,两个 β 值跟默认不太一样:

torch.optim.AdamW(groups, lr=3e-4, betas=(0.9, 0.95), fused=True)

beta2 默认是 0.999,LLM 训练普遍改成 0.95。beta2 控制二阶动量(梯度平方的滑动平均)的窗口长度,0.999 大致是看最近一千步,0.95 大致是看最近二十步。语言模型的梯度分布变化比较快,窗口短一些反应更及时,实践中更稳。

fused=True 把优化器的逐元素运算融合成一个 kernel。参数张量多的时候(我们有一百多个),能省下可观的 kernel 启动开销。只在 CUDA 上可用。

4.2 weight decay 不能一视同仁

这是个容易写错的地方。weight_decay 会把参数往 0 拉,起正则化作用。但不是所有参数都该被拉向 0

RMSNorm 的缩放系数 weight 初始化是 1,含义是「不做缩放」。把它往 0 拉,等于强行削弱归一化的输出,纯属帮倒忙。各种 1 维参数(bias、norm 系数)都是同理。

通行做法是按维度分组:

def make_optimizer(model, args):
decay, no_decay = [], []
for _, p in model.named_parameters():
if not p.requires_grad:
continue
(decay if p.dim() >= 2 else no_decay).append(p)

groups = [
{"params": decay, "weight_decay": args.weight_decay},
{"params": no_decay, "weight_decay": 0.0},
]
return torch.optim.AdamW(groups, lr=args.lr,
betas=(args.beta1, args.beta2),
fused=torch.cuda.is_available())

按维度分而不是按名字分,是因为维度是个更可靠的判据:所有权重矩阵都是 2 维以上,所有 norm 系数和 bias 都是 1 维。按名字匹配容易漏掉改过名的层。

weight decay 把参数往 0 拉 —— 但不是所有参数都该被拉向 0模型全部参数判据:p.dim()dim ≥ 2 → weight_decay = 0.1q/k/v/o_proj、gate/up/down_proj、embedding —— 全是权重矩阵往 0 拉就是正则化,正是 decay 该干的事dim == 1 → weight_decay = 0.037 个 RMSNorm 系数(18 层 × 2 + 最后 1 个)初始化是 1,含义是「不缩放」;拉向 0 等于强行削弱归一化按维度分而不是按名字分:所有权重矩阵都是 2 维以上,所有 norm 系数和 bias 都是 1 维。按名字匹配容易漏掉改过名的层,而漏掉不报错。我们的模型没有 bias(02 篇 2.3 节),所以 no_decay 组里只有那 37 个系数,参数量小到可以忽略 —— 但该分还是要分,因为它影响的是归一化能不�能正常工作。
这是那种「不分也能训出来,但一直有一点点不对」的设置。它不会让 loss 爆炸,只会让归一化层被持续施加一个方向错误的压力。这类问题的共同特征是:没有报错、没有明显症状、只能靠知道原理来避免。

我们的模型没有 bias(02 篇 2.3 节),所以 no_decay 组里只有 37 个 RMSNorm 的系数(18 层 × 2 + 最后 1 个),参数量很少。但该分还是要分。

五、学习率调度

5.1 warmup:前面几百步必须慢

训练一开始,参数是随机初始化的,梯度方向基本没有意义。这时候用大学习率,等于朝着随机方向迈一大步,很容易把模型推到一个很难恢复的区域。

而且 Adam 的二阶动量在最初几步样本太少,估计不准,算出来的步长可能异常大。

warmup 就是让学习率从接近 0 线性升到目标值,给模型一个缓冲。通常取总步数的 1% 到 5%,我们用 2%,也就是 190 步。

5.2 之后 cosine 衰减

warmup 结束后,学习率按余弦曲线降到一个最小值(通常是峰值的 10%):

lr(t)=lrmin+12(1+cos(πttwarmupTtwarmup))(lrmaxlrmin)\text{lr}(t) = \text{lr}_{\min} + \frac{1}{2}\left(1 + \cos\left(\pi \cdot \frac{t - t_{\text{warmup}}}{T - t_{\text{warmup}}}\right)\right)(\text{lr}_{\max} - \text{lr}_{\min})

为什么用余弦而不是线性。余弦的形状是「开头降得慢,中间降得快,结尾又降得慢」。开头慢是因为这时候模型还在快速学习,不该急着降;结尾慢是为了让参数在最后阶段精细收敛。实践中它比线性衰减稳定地好一点。

为什么不降到 0。降到 0 意味着最后一段完全不学习了,浪费算力;而且如果之后想接着训(比如加数据继续训),从 0 起步很别扭。留 10% 是通行做法。

5.3 代码和实际曲线

def lr_at(step, total_steps, warmup_steps, args):
min_lr = args.lr * args.min_lr_ratio
if step < warmup_steps:
return args.lr * (step + 1) / warmup_steps
progress = min(1.0, (step - warmup_steps) / max(1, total_steps - warmup_steps))
return min_lr + 0.5 * (1 + math.cos(math.pi * progress)) * (args.lr - min_lr)

曲线形状:

3e-42.25e-41.5e-47.5e-5002384476871529537训练步数(总 9537 步)3e-5 下限warmup 190 步(2%)线性升到峰值,防止初期梯度方向随机时迈大步cosine 衰减首尾降得慢、中段降得快曲线由 00-pretrain-0.5b/train_schedule.py 的 lr_at() 逐点算出,与实际训练用的调度完全一致。
两段式调度。左端 190 步的窄带是 warmup,之后是余弦衰减到峰值的 10%。不降到 0 有两个原因:末段完全不学习是浪费算力;且后续若要追加数据继续训练,从 0 起步无法平滑衔接。

00-pretrain-0.5b/train_schedule.py 打出来的数值:

学习率曲线
step 进度 lr
0 0.0% 1.579e-06
95 1.0% 1.516e-04
189 2.0% 3.000e-04 ← warmup 结束
190 2.0% 3.000e-04
2384 25.0% 2.649e-04
4768 50.0% 1.693e-04
7152 75.0% 7.110e-05
9536 100.0% 3.000e-05
最低学习率 3.000e-05

第 0 步是 1.579e-6,接近 0;190 步升到 3e-4;然后一路余弦降到 3e-5。

学习率要在每步开头手动写进 optimizer:

lr = lr_at(step, total_steps, warmup_steps, args)
for g in optimizer.param_groups:
g["lr"] = lr

torch.optim.lr_scheduler 也行,但手写这几行更直白,也免得 scheduler 的 step() 和梯度累积的节奏搞混。

5.4 学习率取多大

3e-4 是 0.5B 这个量级的常见取值。经验规律是模型越大,学习率越小:1B 以下常见 3e-4 到 6e-4,7B 量级降到 1e-4 附近,百亿以上更小。

另一条规律是学习率和 batch size 正相关,batch 变大可以适当调高学习率。但这个关系不是线性的,实际还是以试为准。

小模型上调学习率的成本不高,跑几百步看 loss 曲线就知道大概。真拿不准就先用 3e-4,它很少会错得离谱。

六、梯度裁剪

即使有了 warmup,训练过程中偶尔还是会遇到异常样本导致梯度突然变得很大。一次这样的更新就可能把模型带偏。

梯度裁剪的做法是:算出所有梯度拼成的向量的 L2 范数,如果超过阈值,就按比例缩小到阈值。

grad_norm = torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)

阈值取 1.0 是通行值。这个函数会返回裁剪之前的范数,这个返回值很有用,第九节讲怎么用它做监控。

必须在 optimizer.step() 之前、所有梯度累积完成之后调用。放在累积循环里面就变成对每个 micro batch 分别裁剪,语义完全不同。

左:裁剪前后的 grad norm  右:三种该警惕的形态grad norm01.0裁剪阈值 1.0虚线 = 裁剪前的真实范数实线 = 实际用于更新的范数,尖峰被削平异常样本引起的一次尖峰,就足以把模型带偏 —— 裁剪的全部作用就是不让这一步生效① 持续贴着阈值梯度普遍偏大,学习率可能太高② 突然飙到几十几百碰到异常数据,或数值出了问题③ 快速趋近 0模型不学了,学习率太低或断了梯度clip_grad_norm_ 必须在所有累积完成之后、optimizer.step() 之前调用。放进累积循环里,就变成对每个 micro batch 分别裁剪,语义完全不同。它的返回值是裁剪「之前」的范数,也就是上面那条虚线 —— loss 是滞后指标,grad norm 是即时指标,异常往往先出现在这条曲线上。
右边三种形态都不会让训练立刻停下,这正是要主动盯它的原因。第①种最容易被忽略:曲线看着平稳,实际上每一步都在被裁剪,等于学习率被隐式压住了,模型在用一个你没设定的速度学习。

七、完整循环

前面的部件拼起来:

model.train()
t0 = time.time()
for step in range(total_steps):
# 1. 设学习率
lr = lr_at(step, total_steps, warmup_steps, args)
for g in optimizer.param_groups:
g["lr"] = lr

# 2. 梯度累积
for micro in range(args.grad_accum):
x, y = get_batch(train_data, args.micro_batch, args.seq_len, device)
with ctx:
_, loss = model(x, y)
loss = loss / args.grad_accum
loss.backward()

# 3. 裁剪 + 更新 + 清零
grad_norm = torch.nn.utils.clip_grad_norm_(model.parameters(), args.grad_clip)
optimizer.step()
optimizer.zero_grad(set_to_none=True)

# 4. 日志 / 评估 / 存盘
if step % args.log_every == 0:
...
if step % args.eval_every == 0 and step > 0:
vl = evaluate(model, val_data, args, device, ctx)
if step % args.ckpt_every == 0 and step > 0:
save_ckpt(...)

核心就这么多。完整可跑版本在 00-pretrain-0.5b/train.py

顺带说 torch.compile。加一行 model = torch.compile(model) 通常能提速百分之十几到三十,原理是把多个小算子融合成大 kernel,减少显存读写。代价是首次运行要花一两分钟编译。调试阶段建议关掉,因为报错信息会变得很难读。

八、checkpoint

8.1 要存的不只是权重

只存 model.state_dict() 是不够的,那样只能拿来做推理,没法接着训。断点续训至少还要:

要存的东西不存会怎样
模型权重没得恢复
优化器状态Adam 的两个动量丢了,恢复后要重新预热,loss 会跳
当前 step学习率调度对不上,会从头 warmup
模型配置加载时不知道该建多大的模型
随机数状态严格复现需要,不复现可以不存
torch.save({
"model": model.state_dict(),
"optimizer": optimizer.state_dict(),
"step": step,
"model_config": asdict(cfg),
"train_args": asdict(args),
"rng": torch.get_rng_state(),
"np_rng": np.random.get_state(),
}, path)

优化器状态往往比权重还大。 AdamW 每个参数存两个 fp32 动量,就是 8 字节每参数,而 bf16 权重只有 2 字节。0.5B 模型的 checkpoint 大约是 1 GB 权重加 4 GB 优化器状态。存盘要留够磁盘空间。

8.2 必须原子写

这是租卡场景下的一个真实教训。直接 torch.save(path) 存一个 5 GB 的 checkpoint 要几十秒,如果这期间实例被回收或者进程被杀,会留下一个写了一半的文件,而它已经把上一个好的 checkpoint 覆盖了。等于两个都没了。

对策是先写临时文件,写完再 rename:

def save_ckpt(path, ...):
tmp = path + ".tmp"
torch.save({...}, tmp)
os.replace(tmp, path) # 同一文件系统上,rename 是原子操作

os.replace 在同一个文件系统内是原子的,要么完全成功,要么原文件保持不变,不存在中间状态。

8.3 存盘频率

太频繁浪费时间(每次几十秒),太稀疏丢的进度多。按 04 篇的配置,一步大约 6 秒,500 步存一次就是每 50 分钟一次,最多丢 50 分钟。这个粒度比较合适。

另外记得存到持久化云盘,不要存在实例的本地盘上。这点 00 篇 3.4 节强调过。

checkpoint 要存的五样东西(只存权重只能拿去推理,接不上训练)模型权重不存就没得恢复优化器状态Adam 双动量丢了,恢复后 loss 会跳当前 step学习率调度对不上,会从头 warmup模型配置加载时不知道该建多大的模型随机数状态严格复现才需要,不复现可以不存写 5 GB 要几十秒,而实例可能在任何一秒被回收❌ 直接写torch.save(path) 直接覆盖旧文件,写 5 GB 持续几十秒中途被回收 → 半个文件,旧的也没了✅ 原子写先写 path + ".tmp",旧文件全程原封不动os.replace任何时刻被回收,旧 checkpoint 都完好os.replace 在同一个文件系统内是原子的:要么完全成功,要么原文件保持不变,不存在中间状态。跨文件系统就不成立,临时文件必须和目标放在同一个盘。存盘频率:一步约 6 秒,500 步存一次约等于每 50 分钟,最多丢 50 分钟进度。存到持久化云盘,不要存实例本地盘。
「两个都没了」是这张图里最贵的一格:直接覆盖写的失败模式不是丢掉这一次存盘,而是把上一个能用的 checkpoint 一起毁掉。三行代码换来的保证是——无论进程在哪一微秒被杀,磁盘上永远有一个完整可加载的文件。

九、训练时该盯什么

9.1 loss

最基本的。开头应该是 10.37(02 篇 8.3 节),然后快速往下掉。

TinyStories 这种简单语料,几百步就能掉到 3 以下。FineWeb-Edu 会慢很多,但趋势应该一直向下。

同时看训练 loss 和验证 loss。两者持续接近说明健康;验证 loss 开始上升而训练 loss 继续降,是过拟合。不过预训练阶段数据量远大于参数量,一般看不到过拟合。

9.2 grad norm

就是第六节 clip_grad_norm_ 的返回值,这是个被低估的指标。

正常情况下它应该在一个稳定范围内小幅波动。值得警惕的是:

  • 持续贴着裁剪阈值(一直是 1.0):说明梯度普遍偏大,学习率可能太高
  • 突然飙到几十几百:碰到异常数据,或者数值出了问题
  • 快速趋近 0:模型不学了,可能学习率太低或者哪里断了梯度

loss 曲线是滞后指标,grad norm 往往能更早暴露问题。

9.3 MFU 怎么算

MFU(Model FLOPs Utilization)衡量硬件用得有多充分,定义是实际有效算力除以理论峰值:

MFU=6×N×tokens/s峰值 FLOPS×卡数\text{MFU} = \frac{6 \times N \times \text{tokens/s}}{\text{峰值 FLOPS} \times \text{卡数}}

分子是 00 篇那个 6ND6ND 的每秒版本。代码里:

tps = tokens_per_step * args.log_every / dt
mfu = 6 * model.num_params() * tps / 312e12

00-pretrain-0.5b/train_schedule.py 打出了耗时和 MFU 的对应关系:

耗时与 MFU 的对应关系
耗时(小时) tokens/s MFU
8 347,222 83.8%
12 231,481 55.9%
17 163,399 39.5%
24 115,741 27.9%
36 77,160 18.6%
67 41,459 10.0%

00 篇按 MFU = 0.40 估的 17 小时,正好对应表里第三行的 39.5%,两边自洽。

同一份工作量,跑多久完全由 MFU 决定MFU084%30% 警戒线8 h17 h36 h67 h17 小时 · 39.5% · 163,399 tokens/s ← 00 篇按 0.40 估的就是这里MFU 低于 30% 该查什么(按概率排序)① 数据加载跟不上CPU 或磁盘瓶颈,GPU 在等数据② micro_batch 太小GPU 利用不满,kernel 启动开销占比高③ 没开 torch.compile④ 多卡通信占比过高 —— 04 篇细说MFU = 6N × tokens/s ÷(峰值 FLOPS × 卡数),分子就是 00 篇那个 6ND 的每秒版本。两边用的是同一套账,所以能互相验算。
这条曲线的形状值得记住:MFU 从 40% 掉到 28%,耗时从 17 小时涨到 24 小时;掉到 10%,就是 67 小时。也就是说排查性能问题的收益是非线性的,越靠近低 MFU 区间,每提升一个百分点省下的时间越多 —— 这也是为什么 30% 是个值得停下来查的界线。

实际跑起来,MFU 低于 30% 就该查瓶颈。常见原因按概率排序:数据加载跟不上(CPU 或磁盘瓶颈)、micro_batch 太小导致 GPU 利用不满、没开 torch.compile、多卡时通信占比过高(04 篇细说)。

9.4 tokens/s

比 MFU 更直观,用来估还要跑多久。拿它除一下剩余 token 数就是剩余时间。

十、先在 TinyStories 上跑通

跟 01 篇一样,别一上来就开四卡烧 FineWeb-Edu。先用小配置在单卡上把整条链路走通。

小配置这样改:

Config(vocab_size=32000, d_model=256, n_layers=4,
n_heads=4, n_kv_heads=2, ffn_dim=688, max_seq_len=512)

配套的训练参数把 total_tokens 调到几千万,micro_batch 可以开大些,compile=False 方便看报错。

要确认的事:

  • 第 0 步 loss 在 10.37 附近
  • loss 确实在下降,几百步内能明显低于 10
  • grad norm 在合理范围,没有一直贴着 1.0
  • 学习率曲线跟 5.3 节打印的一致
  • 存一次 checkpoint,杀掉进程,能从 checkpoint 恢复且 loss 接续得上
  • 恢复后学习率跟中断前对得上
  • MFU 有数(小模型 MFU 天然低,看不到 40%,正常)

第 5 条最值得认真测。 断点续训是那种「以为写对了,真出事时才发现不行」的功能。主动杀一次进程验证,比事后后悔便宜太多。

十一、常见问题

现象大概率是什么原因
loss 一开始就不是 10.37模型有问题,回 02 篇八节验收
loss 降到某处变 NaN忘了除 grad_accum;或学习率太高;或没做梯度裁剪
loss 完全不降学习率设成 0 了;或忘了 optimizer.step();或 warmup 步数写成了总步数
loss 降得异常快因果掩码失效,模型在抄答案,回 02 篇 8.4 节
OOM调小 micro_batch,相应调大 grad_accum 保持总 batch 不变
MFU 只有百分之十几数据加载瓶颈,检查 pin_memory;或 micro_batch 太小
grad norm 一直是 1.0学习率偏高,或数据里有异常样本
恢复 checkpoint 后 loss 跳一下优化器状态没存,只存了权重
恢复后学习率不对step 没存,调度从头开始了
torch.compile 报一堆看不懂的错先设 compile=False 定位真正的问题

这一篇的产出要能被杀死

判断 train.py 写完的标准不是「能跑」,是「跑到一半 kill -9,重启之后能接着跑,loss 曲线看不出接缝」。租卡环境下实例随时可能没,这个能力比什么都实在。